Conv2DBackpropInputFusion
计算二维卷积对输入的反向梯度(input gradient)。
dy、dx 为 NHWC;w 为卷积核布局(见下);支持普通卷积、Depthwise,以及满足条件时的 1x1 优化路径。
对组 \(g\)、输出位置 \((o_h, o_w)\)、核偏移 \((k_h, k_w)\),令 \(C_o' = C_{out}/G\),\(C_i' = C_{in}/G\)。对应输入坐标为:
当 \((i_h, i_w)\) 落在输入有效范围内时累加,否则该项为 0。
实现上:普通路径对 dy 与 w 做 GEMM(GEneral Matrix Multiply,通用矩阵乘),
再经 col2im(column to image,列展开还原成特征图)累加到 dx;
Depthwise 与 1x1(无 pad、dilation/stride 为 1、group 为 1)走专用路径。
- 输入:
dy - 输出梯度地址,形状 [batch, out_h, out_w, out_channel](NHWC)
w - 权重地址,形状 [out_channel, in_channel/group, kernel_h, kernel_w] (OIHW,不是 OHWI / NHWC)
conv_param - ConvParameter 结构体地址,字段见下
core_mask - 核掩码(仅共享存储版本使用)
ConvParameter 定义:
1typedef struct ConvParameter {
2 void* workspace_; // 工作区指针;实际分配字节数见下方警告(常大于 buffer_size_)
3 int output_batch_; // dy 的 batch
4 int input_batch_; // dx 的 batch(通常与 output_batch_ 相同)
5 int input_h_; // 输入高(dx)
6 int input_w_; // 输入宽(dx)
7 int output_h_; // 输出梯度高(dy)
8 int output_w_; // 输出梯度宽(dy)
9 int input_channel_; // 输入通道(dx)
10 int output_channel_; // 输出通道(dy)
11 int kernel_h_; // 卷积核高
12 int kernel_w_; // 卷积核宽
13 int group_; // 组数,1 为普通卷积
14 int pad_l_; // 左填充
15 int pad_u_; // 上填充
16 int dilation_h_; // 高方向膨胀
17 int dilation_w_; // 宽方向膨胀
18 int stride_h_; // 高方向步长
19 int stride_w_; // 宽方向步长
20 int buffer_size_; // 分块预算,单位字节(非元素个数);_s/_p 都读此字段。见下方警告。
21 int pad_d_; // 下填充
22 int pad_r_; // 右填充
23 int nweights_; // 权重 w 的元素总数
24} ConvParameter;
- 输出:
dx - 输入梯度地址,形状
[batch, in_h, in_w, in_channel]; 调用前通常先清零,再由算子累加写入
- 支持平台:
FT78NEMT7004
警告
MT7004的workspace_ / buffer_size_ 分配算法(与前向 conv2d / transpose / backprop_filter 都不同):
权重布局是 OIHW [C_out, C_in/G, K_h, K_w],不是前向常见的 OHWI。传错布局结果静默错误。
必须填对 nweights_ = out_channel × (in_channel / group) × kernel_h × kernel_w, 且与 w 实际元素数一致;组卷积按 nweights_/group 切各组权重起始。
buffer_size_ 是字节(不是元素个数)。内部用 buffer_size_/2 估半区;部分汇编按 float 个数做指针算术,因此 workspace_ 须 ≥ buffer_size_ × sizeof(float)。字段与分配字节可以不相等。
共享 _s 与私有 _p 都会读 buffer_size_ / workspace_ (不像前向 conv _s 忽略字段)。无 transpose 那种 packed 权重 + core_num×2048 的额外 DDR scratch 规则; 也无 filter 那种 core_num×dw partial 规则。常用:
buffer_size_ = max(2048, 2 * Kh * Kw * (Cin / G) * sizeof(dtype)); workspace_bytes = buffer_size_ * sizeof(float);
测试常用 buffer_size_ = 0x10000。
备注
FT78NE 支持 fp32
MT7004 支持 fp16、fp32
Depthwise 条件:output_channel、group、input_channel 三者相等,且 dilation 为 1
1x1 优化条件:Kh=Kw=1,四向 pad 为 0,dilation/stride 为 1,group 为 1
共享存储版本:
-
void hp_conv2d_backprop_input_s(float16 *dy, float16 *w, float16 *dx, ConvParameter *conv_param, int core_mask)
-
void fp_conv2d_backprop_input_s(float *dy, float *w, float *dx, ConvParameter *conv_param, int core_mask)
C调用示例:
1// MT7004 示例(共享存储多核,DDR 地址)
2void TestConvBackpropInputSMCFp32(int core_mask) {
3 int core_id = get_core_id();
4 int logic_core_id = GetLogicCoreId(core_mask, core_id);
5 int core_num = GetCoreNum(core_mask);
6 float *dy = (float *)0x86000000;
7 float *w = (float *)0x87000000; // OIHW: [16, 16, 1, 1]
8 float *dx = (float *)0x88000000;
9 ConvParameter *param = (ConvParameter *)0x89000000;
10 float *workspace = (float *)0x8A000000;
11 if (logic_core_id == 0) {
12 param->workspace_ = workspace;
13 // 字段填字节预算;实际 malloc 建议 >= buffer_size_ * sizeof(float)
14 param->buffer_size_ = 0x10000;
15 param->input_batch_ = 1;
16 param->output_batch_ = 1;
17 param->input_h_ = 4;
18 param->input_w_ = 4;
19 param->input_channel_ = 16;
20 param->output_h_ = 4;
21 param->output_w_ = 4;
22 param->output_channel_ = 16;
23 param->kernel_h_ = 1;
24 param->kernel_w_ = 1;
25 param->group_ = 1;
26 param->pad_l_ = 0;
27 param->pad_u_ = 0;
28 param->pad_d_ = 0;
29 param->pad_r_ = 0;
30 param->dilation_h_ = 1;
31 param->dilation_w_ = 1;
32 param->stride_h_ = 1;
33 param->stride_w_ = 1;
34 param->nweights_ = 16 * 16 * 1 * 1;
35 }
36 sys_bar(0, core_num);
37 fp_conv2d_backprop_input_s(dy, w, dx, param, core_mask);
38}
39
40void main() {
41 int core_mask = 0b1111;
42 TestConvBackpropInputSMCFp32(core_mask);
43}
私有存储版本:
-
void hp_conv2d_backprop_input_p(float16 *dy, float16 *w, float16 *dx, ConvParameter *conv_param)
-
void fp_conv2d_backprop_input_p(float *dy, float *w, float *dx, ConvParameter *conv_param)
C调用示例:
1// MT7004 示例(私有存储单核,AM 地址)
2void TestConvBackpropInputAMFp32(void) {
3 float *dy = (float *)0x10000000;
4 float *w = (float *)0x10010000; // OIHW: [16, 16, 1, 1]
5 float *dx = (float *)0x10020000;
6 ConvParameter *param = (ConvParameter *)0x10030000;
7 float *workspace = (float *)0x10040000;
8 param->workspace_ = workspace;
9 // 字段填字节预算;workspace 建议至少 buffer_size_ * sizeof(float)
10 param->buffer_size_ = 0x10000;
11 param->input_batch_ = 1;
12 param->output_batch_ = 1;
13 param->input_h_ = 4;
14 param->input_w_ = 4;
15 param->input_channel_ = 16;
16 param->output_h_ = 4;
17 param->output_w_ = 4;
18 param->output_channel_ = 16;
19 param->kernel_h_ = 1;
20 param->kernel_w_ = 1;
21 param->group_ = 1;
22 param->pad_l_ = 0;
23 param->pad_u_ = 0;
24 param->pad_d_ = 0;
25 param->pad_r_ = 0;
26 param->dilation_h_ = 1;
27 param->dilation_w_ = 1;
28 param->stride_h_ = 1;
29 param->stride_w_ = 1;
30 param->nweights_ = 16 * 16 * 1 * 1;
31 fp_conv2d_backprop_input_p(dy, w, dx, param);
32}
33
34void main() {
35 TestConvBackpropInputAMFp32();
36}